請允許我再水一篇XD 一直水這個主題XD
我明天,我希望啦,從明天開始我會水其他一個主題了。
然後我之後我會花個一兩三天,認認真真地把這東西寫一個我覺得好的技術 blog。但這幾天就請讓我再水一下吧。
--
昨天我們做完 DoubleDouble。
DoubleDouble 的做法是把一個數拆成:
value = hi + lo
這樣可以從 FP64 的 53-bit 有效精度,往上推到大約 106-bit。
今天再往前推一步:做 FP128。
這裡的 FP128 指的是我們程式裡使用的 Float128 路線:113-bit significand,加上 exponent 和 sign。它的目的很直接:
不要只保住 DoubleDouble 的 106-bit
而是把有效精度補到 113-bit
也就是從昨天的 DoubleDouble,再多補 7 bit。
昨天 DoubleDouble 已經解決一個問題:lo 不能被 collapse 成 double。
今天 FP128 要解的是下一個問題:
DoubleDouble 只有約 106-bit
FP128 需要 113-bit
所以不能拿 DoubleDouble 改名。
今天要補的是這幾個地方:
1. 用新的 Float128 表示 113-bit significand
2. modulus count 從 DoubleDouble 的 29/30 顆往上加
3. residue serialization 不能經過 double,也不能經過 DoubleDouble
4. CRT reconstruction 的 signed range 要多留 1 bit
5. 驗證要打到 2^-110 這種 DoubleDouble/FP64 容易看不到的位置
這些補完,FP128 path 才是真的接進同一套 scheme。
這次新增的格式不是:
double hi;
double lo;
而是直接存 significand:
struct Float128 {
uint64_t mant_lo;
uint64_t mant_hi;
int32_t expo;
int32_t sign;
};
可以把它想成:
|value| = M * 2^expo
其中 M 是 113-bit 的整數 significand。
normalized 之後:
2^112 <= M < 2^113
這個表示法對我們很方便,因為後面本來就要把數字丟進整數和 modulo 世界。
所以 FP128 不是先變成 double,也不是先變成 DoubleDouble。
它直接走:
Float128(M, expo, sign)
-> scaled integer
-> residue
這是今天最重要的差異。
昨天 DoubleDouble 的有效精度大約是 106-bit。
今天 FP128 要用 113-bit。
GEMM 裡每個輸出元素是 K 個乘積累加,所以 CRT 空間要看:
required_crt_bits = 2 * input_bits + ceil(log2(K)) + 1
前面三項很直覺:
input_bits 乘 input_bits -> 2 * input_bits
K 項相加 -> ceil(log2(K))
最後這個 +1 是 signed CRT range。
我們重建時不是只看 [0, P),而是要把結果解釋成:
(-P/2, P/2)
如果少留這一個 bit,某些正數結果可能會落到 P/2 後面,最後被當成負數。
所以 FP128、K=4096 時:
2 * 113 + ceil(log2(4096)) + 1
= 226 + 12 + 1
= 239 bits
這就是今天 modulus count 必須上升的原因。
目前這組 INT8 modulus 表有 49 顆可以用。
全部乘起來大約是:
341.88 bits
但 FP128 不需要全拿。
重點是前綴乘積要跨過 required bits。
以 K=4096 來看:
| modulus count | CRT bits | 結果 |
|---|---|---|
| 30 | 約 226.9 | 不夠 |
| 31 | 約 233.7 | 不夠 |
| 32 | 約 240.4 | 夠 |
所以 FP128 path 會選:
num_moduli = 32
如果硬指定 30 顆去跑 K=4096,程式會直接拒絕:
Insufficient CRT bits for FP128
這裡不是保守。
這是正確性條件。
30 顆對 DoubleDouble 可能夠,對 FP128 不夠。
昨天 DoubleDouble 要先把 hi/lo 合成 scaled integer words。
今天 Float128 更直接,因為它本身就是:
M * 2^expo
所以序列化就是:
scaled = M * 2^(expo + scale)
程式裡對應到:
Float128::to_scaled(UInt384* mag, bool* neg, int scale)
to_scaled 做完後,我們拿到:
UInt384 mag
bool neg
接著對每個 modulus 做:
residue = mag % p
如果是負數,就轉成該 modulus 底下的負數 residue。
最後再轉成 balanced int8:
if residue > p / 2:
residue -= p
所以每個 Float128 element 的資料路徑是:
Float128
-> UInt384 scaled integer
-> 32 個 residues
-> 32 張 int8 residue matrix
這裡不能偷走:
Float128 -> double -> residue
Float128 -> DoubleDouble -> residue
因為第 107 到 113 bit 會在這裡掉光。
到了 GEMM 這段,形式跟前面一樣:
C_mod[p] = A_mod[p] * B_mod[p]
差別是這次 modulus count 變成 32。
也就是:
p0, p1, ..., p31
每一個 modulus 都有一組 int8 GEMM。
輸入:
A_mod[p]: int8
B_mod[p]: int8
輸出:
C_mod[p]: int32
這一步不負責 FP128 round,也不負責最後輸出格式。
它只負責把每個 modulus 底下的矩陣乘法算對。
高精度是在下一步 CRT reconstruction 才拼回來。
32 個 C_mod[p] 算完後,對每個輸出元素做 CRT。
重建用:
M_total = product(p_i)
M_i = M_total / p_i
inv_i = inverse(M_i mod p_i)
然後:
result = sum C_mod[i] * inv_i * M_i
result = result mod M_total
因為 FP128 K=4096 需要 239 bits,UInt384 夠用。
重建後要先判斷 signed range:
result > M_total / 2 -> negative
最後回到 Float128:
Float128::from_scaled(abs_result, sign, scale)
這一步會把寬整數重新 normalize 成 113-bit significand。
也就是回到:
mant_hi
mant_lo
expo
sign
這樣輸出才還是 FP128 path,不會掉回 double。
如果只是測普通 random matrix,沒有什麼說服力。
今天要測的是 FP64 和 DoubleDouble-from-double 都看不到,但 Float128 看得到的位置。
我們用:
2^-110
測 cancellation。
設定:
A = +/-1 + 2^-110
B = 1
K = 128
如果先 collapse 成 double,2^-110 直接不見。
所以 FP64 看到的是 +1 和 -1,最後相消:
FP64 result = 0
但 Float128 會保留每一項的 2^-110。
K=128 時應該留下:
128 * 2^-110 = 9.86076e-32
實測:
FP128 result: 9.86076e-32
FP64 collapsed result: 0
DD-from-double result: 0
oracle: 9.86076e-32
這個測試很直接。
它證明 2^-110 沒有在 input serialization、residue GEMM、CRT reconstruction 裡被吃掉。
目前 128x128x128 跑了三組。
| case | FP128 max_rel | FP64 / DD-from-double |
|---|---|---|
cancellation +/-1 + 2^-110 |
0 | 1.0 |
constant 1 + 2^-110 |
0 | double 看起來都是 128 |
varied + 2^-110 / 2^-108 |
約 1.23e-28 | 約 7.92e-12 / 8.80e-12 |
還有一個 scalar check:
1
1 + 2^-110
轉成 double 會是同一個值。
但在 Float128 裡,significand 不同。
這個 check 用來確認 input format 一開始就沒有掉 precision。
把今天的流程合起來就是:
Float128 input
-> to_scaled(UInt384)
-> UInt384 % p_i
-> balanced int8 residue matrix
-> 32 個 int8 GEMM
-> 32 個 int32 C_mod
-> UInt384 CRT reconstruction
-> Float128::from_scaled
-> Float128 output
這條 path 已經跑通。
它和昨天 DoubleDouble 的差別不是概念變了,而是 precision target 變了:
106-bit -> 113-bit
30 moduli 左右 -> 32 moduli
DoubleDouble hi/lo -> Float128 significand/exponent
DoubleDouble output -> Float128 output
這版還是 correctness path。
從昨天 DoubleDouble 的 profile 已經看得出來,真正大的時間不在中間那段 int8 GEMM,而是在前後:
高精度 input -> residue
CRT reconstruction -> 高精度 output
host/device round trip
FP128 會更明顯,因為它有 32 個 modulus。
所以下一步如果要加速,不是只改 GEMM kernel。
要把這些東西搬到 GPU 裡:
Float128 to_scaled
tile-local residue generation
32-modulus fixed schedule
accumulator-side mod reduction
partial CRT epilogue
GPU-side limb reconstruction
也就是把現在的 host correctness path 壓進真正 fused 的 GPU pipeline。
昨天 DoubleDouble 證明 106-bit path 可以跑。
今天 FP128 把 target 推到 113-bit。
關鍵改動是:
Float128 value format
signed CRT +1 bit
32 個 modulus
UInt384 to_scaled / from_scaled
2^-110 cancellation 驗證
所以 Day 14 不是規劃。
這版已經把 FP128 的 113-bit correctness path 接進來了。
下一步才是把它從 correctness path,推到 GPU fused path。
為了讓AI文有一點人味 於是我會寫一些 題外mur mur
嗯,寫這個要不斷堅持下去,我真的覺得太麻煩太困難了。
我覺得齁
人阿 還是要很老實地去面對
交踏實地的面對發生的各種事情。
你知道就是,筆者其實只是一個一般的碩士生,然後快要畢業,還沒畢業這樣的一個碩士生。
那其實我一開始的想法是
哇,這年頭真的是 junior找工很難找啊,
AI到底在幹嘛阿,弄得我們找工直接夢魘模式。
然後你知道,我一開始學這個也是想說
GPU,大家都那麼需要 AI
GPU弄弄應該找工 輕輕鬆鬆吧?
所以我也去報了那種開源的講題的議程
(結果還報了錯,聽說人家另外一個議程報一報還有拿硬體的,哇靠,早知道我也報那個。)
反正就是弄弄之後,自己其實也,我也不覺得真的有怎麼樣吧,
雖然其實我也沒開始投。然後就是我也不知道。
我覺得人還是要認份啦,該刷題還是要刷題啦。
我最近也開始在刷題了。
對,大家一起努力吧,雖然我也不知道現在刷題還能不能夠有好的工作,
但是
一起努力,一起看齊。
已閱讀6913字。